#!/usr/bin/env python3
"""Stage 3 Task 0: build C-###_query_info_legal_facts.json from claim descriptions."""

from __future__ import annotations

import itertools
import json
import re
import sys
from pathlib import Path
from typing import Any

import httpx

INPUT_DIR = Path("Stage_2_Results/Task_D")
OUTPUT_DIR = Path("Stage_3_Results/Task_0")

# Fallback claim ids for code-executor environments where the workspace folder
# is not directly visible and files must be read through localdocs.
FALLBACK_CLAIM_IDS = [
    "C-001",
    "C-002",
    "C-003",
    "C-004",
    "C-005",
    "C-006",
    "C-007",
    "C-008",
    "C-009",
    "C-010",
]

TARGET_FIELDS = [
    "claim_id",
    "case_kind",
    "claim_title",
    "plaintiffs",
    "defendants",
    "claim_statement",
    "relief_summary",
    "cause_summary",
]

INPUT_NAME_RE = re.compile(r"^(C-\d{3})_claim_description\.json$")

LOCALDOCS_URL = "http://mcp-localdocs:8012/mcp"
MCP_HEADERS = {
    "Content-Type": "application/json",
    "Accept": "application/json, text/event-stream",
}
MSG_ID_COUNTER = itertools.count(10)


def log(message: str) -> None:
    print(message, file=sys.stderr)


def next_msg_id() -> int:
    return next(MSG_ID_COUNTER)


def _parse_mcp(text: str) -> Any:
    for line in text.strip().split("\n"):
        if line.startswith("data: "):
            try:
                return json.loads(line[6:])
            except Exception:
                return None
    try:
        return json.loads(text)
    except Exception:
        return None


def _init_mcp_session(client: httpx.Client) -> None:
    response = client.post(
        LOCALDOCS_URL,
        json={
            "jsonrpc": "2.0",
            "id": 1,
            "method": "initialize",
            "params": {
                "protocolVersion": "2025-03-26",
                "capabilities": {},
                "clientInfo": {"name": "stage3-task0", "version": "1.0"},
            },
        },
        headers=MCP_HEADERS,
    )
    response.raise_for_status()
    session_id = response.headers.get("mcp-session-id")
    if session_id:
        MCP_HEADERS["mcp-session-id"] = session_id
    client.post(
        LOCALDOCS_URL,
        json={"jsonrpc": "2.0", "method": "notifications/initialized"},
        headers=MCP_HEADERS,
    ).raise_for_status()


def _call_tool(client: httpx.Client, name: str, arguments: dict[str, Any], msg_id: int) -> dict[str, Any]:
    response = client.post(
        LOCALDOCS_URL,
        json={
            "jsonrpc": "2.0",
            "id": msg_id,
            "method": "tools/call",
            "params": {"name": name, "arguments": arguments},
        },
        headers=MCP_HEADERS,
    )
    response.raise_for_status()
    parsed = _parse_mcp(response.text)
    if not parsed or "result" not in parsed:
        raise RuntimeError(f"MCP {name} failed: {response.text[:300]}")
    return parsed


def _unwrap_json(parsed: dict[str, Any], doc_name: str) -> Any:
    text = (parsed["result"].get("content") or [{}])[0].get("text", "")
    try:
        outer = json.loads(text)
    except Exception as exc:
        raise RuntimeError(f"Invalid JSON wrapper for {doc_name}: {text[:200]}") from exc

    if isinstance(outer, dict) and "results" in outer:
        results = outer.get("results") or []
        if not results:
            raise RuntimeError(f"No results for {doc_name}")
        inner = results[0].get("content", "")
        if not inner:
            raise RuntimeError(f"Empty content for {doc_name}")
        return json.loads(inner) if isinstance(inner, str) else inner

    return outer


def read_json_via_mcp(client: httpx.Client, doc_name: str) -> Any:
    parsed = _call_tool(client, "read_docs", {"doc_names": [doc_name]}, next_msg_id())
    return _unwrap_json(parsed, doc_name)


def write_text_via_mcp(client: httpx.Client, path: str, content: str) -> None:
    _call_tool(
        client,
        "write_file",
        {"path": path, "content": content, "overwrite": True},
        next_msg_id(),
    )


def build_output_name(claim_id: str) -> str:
    return f"{claim_id}_query_info_legal_facts.json"


def extract_query_info(payload: dict[str, Any], fallback_claim_id: str) -> dict[str, Any]:
    extracted = {field: payload.get(field) for field in TARGET_FIELDS}
    if not extracted["claim_id"]:
        extracted["claim_id"] = fallback_claim_id
    return extracted


def process_local_files() -> list[str]:
    if not INPUT_DIR.exists():
        return []

    OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
    written_files: list[str] = []

    for input_path in sorted(INPUT_DIR.glob("C-*_claim_description.json")):
        match = INPUT_NAME_RE.match(input_path.name)
        if not match:
            continue

        claim_id = match.group(1)
        log(f"[Stage3-Task0] Reading local file: {input_path}")
        with input_path.open("r", encoding="utf-8") as handle:
            claim_description = json.load(handle)

        output_payload = extract_query_info(claim_description, claim_id)
        output_path = OUTPUT_DIR / build_output_name(output_payload["claim_id"])
        output_path.write_text(
            json.dumps(output_payload, ensure_ascii=False, indent=2),
            encoding="utf-8",
        )
        written_files.append(output_path.as_posix())
        log(f"[Stage3-Task0] Written local file: {output_path}")

    return written_files


def process_via_mcp() -> list[str]:
    client = httpx.Client(timeout=60)
    _init_mcp_session(client)

    written_files: list[str] = []
    try:
        for claim_id in FALLBACK_CLAIM_IDS:
            input_name = f"{INPUT_DIR.as_posix()}/{claim_id}_claim_description.json"
            output_name = f"{OUTPUT_DIR.as_posix()}/{build_output_name(claim_id)}"

            log(f"[Stage3-Task0] Reading MCP doc: {input_name}")
            claim_description = read_json_via_mcp(client, input_name)
            output_payload = extract_query_info(claim_description, claim_id)

            write_text_via_mcp(
                client,
                output_name,
                json.dumps(output_payload, ensure_ascii=False, indent=2),
            )
            written_files.append(output_name)
            log(f"[Stage3-Task0] Written MCP doc: {output_name}")
    finally:
        client.close()

    return written_files


def main() -> None:
    written_files = process_local_files()
    mode = "local-fs"

    if not written_files:
        log("[Stage3-Task0] Local input directory not available; falling back to localdocs MCP.")
        written_files = process_via_mcp()
        mode = "mcp-localdocs"

    if not written_files:
        raise RuntimeError("No C-###_query_info_legal_facts.json files were generated.")

    summary = {
        "status": "ok",
        "mode": mode,
        "processed_count": len(written_files),
        "output_dir": OUTPUT_DIR.as_posix(),
        "written_files": written_files,
    }
    print(json.dumps(summary, ensure_ascii=False))


if __name__ == "__main__":
    main()
